10. 理解 CausalLM
CausalLM 是什么
很多模型结构会写:
LlamaForCausalLM
QwenForCausalLM
GPTForCausalLM
什么是 Causal LM
Causal 的翻译是:
因果的
意思:
只能看左边
不能看右边
例如:
我 爱 吃 ?
预测:
火锅
模型只能看:
我 爱 吃
不能偷看:
火锅
这就是:
Causal Mask
Attention矩阵:
我 爱 吃 火锅
我 ✓
爱 ✓ ✓
吃 ✓ ✓ ✓
火锅 ✓ ✓ ✓ ✓
右上角全部屏蔽:
我 爱 吃 火锅
我 ✓ ✗ ✗ ✗
爱 ✓ ✓ ✗ ✗
吃 ✓ ✓ ✓ ✗
火锅 ✓ ✓ ✓ ✓
这就是自回归训练。
一、CausalLM 整体架构
先看最外层:
class MiniMindForCausalLM(
PreTrainedModel,
GenerationMixin
)
继承关系:
classDiagram
class PreTrainedModel{
+save_pretrained()
+from_pretrained()
+state_dict()
+load_state_dict()
}
class GenerationMixin{
+generate()
+beam_search()
+sample()
+top_k()
+top_p()
}
class MiniMindForCausalLM{
+forward()
+generate()
}
PreTrainedModel <|-- MiniMindForCausalLM
GenerationMixin <|-- MiniMindForCausalLMPreTrainedModel负责什么
它不是模型。
它只是 HuggingFace 的壳。
提供:
save_pretrained()
from_pretrained()
push_to_hub()
例如:
model.save_pretrained("./ckpt")
保存:
config.json
model.safetensors
加载:
MiniMindForCausalLM.from_pretrained(...)
GenerationMixin负责什么
负责推理生成。
里面封装:
generate()
beam_search()
sample()
所以:
model.generate(...)
其实来自:
GenerationMixin
二、MiniMindForCausalLM
代码:
self.model = MiniMindModel(...)
self.lm_head = nn.Linear(...)
结构:
flowchart TD A[Input IDs] A --> B[MiniMindModel] B --> C[Hidden States] C --> D[LM Head] D --> E[Logits]
本质:
Transformer负责提取特征
LM Head负责预测下一个Token
三、MiniMindModel
源码:
self.embed_tokens
self.layers
self.norm
对应:
flowchart TD A[Input IDs] A --> B[Embedding] B --> C[Block1] C --> D[Block2] D --> E[Block3] E --> F[...] F --> G[BlockN] G --> H[Final RMSNorm] H --> I[Hidden States]
这就是模型主体。
四、Transformer Block
源码:
class MiniMindBlock
这是最核心部分。
Block结构图
flowchart TD X[Input x] X --> N1[RMSNorm] N1 --> ATTN[Attention] ATTN --> ADD1[Add Residual] X ----> ADD1 ADD1 --> N2[RMSNorm] N2 --> FFN[FeedForward] FFN --> ADD2[Add Residual] ADD1 ----> ADD2 ADD2 --> Y[Output]
对应代码:
residual = hidden_states
hidden_states = Attention(
RMSNorm(hidden_states)
)
hidden_states += residual
即:
第二部分:
hidden_states =
hidden_states +
MLP(
RMSNorm(hidden_states)
)
即:
完整公式:
五、Attention 在 Block 中的位置
展开:
flowchart TD X[Input] X --> Q[WQ] X --> K[WK] X --> V[WV] Q --> SCORE[QK^T] K --> SCORE SCORE --> MASK[Causal Mask] MASK --> SOFTMAX SOFTMAX --> MUL[×V] V --> MUL MUL --> WO[WO] WO --> OUT
这就是:
self.self_attn(...)
内部做的事情。
六、为什么要Residual
假设32层:
Layer1
Layer2
Layer3
...
Layer32
如果没有残差:
梯度越来越小
训练会崩。
所以:
输入直接绕过去
输出 += 输入
即:
y = f(x) + x
ResNet思想。
七、为什么Norm放前面
旧GPT:
Attention
↓
Norm
Post-Norm
现代:
Norm
↓
Attention
Pre-Norm
即:
Attention(
RMSNorm(x)
)
优点:
深层训练更稳定
梯度传播更容易
Llama/Qwen都这样。
八、整个Model Forward流程
源码:
hidden_states =
embed_tokens(input_ids)
得到:
[B,L]
↓
[B,L,H]
例如:
[1,128]
↓
[1,128,768]
然后进入:
for layer in self.layers:
即:
flowchart LR X --> B1[Block1] B1 --> B2[Block2] B2 --> B3[Block3] B3 --> B4[...] B4 --> BN[BlockN]
最后:
hidden_states = self.norm(hidden_states)
最终输出:
[B,L,H]
九、LM Head是什么
源码:
self.lm_head
=
Linear(
hidden_size,
vocab_size
)
假设:
hidden_size=768
vocab_size=50000
矩阵:
768 × 50000
作用:
flowchart LR A[768维特征] A --> B[LM Head] B --> C[50000个词概率]
例如:
今天天气很
↓
好 0.82
差 0.05
冷 0.04
...
十、Weight Tying
代码:
self.model.embed_tokens.weight
=
self.lm_head.weight
即:
Embedding矩阵
和
LM Head矩阵
共享参数
图:
flowchart TD TOKENS TOKENS --> EMB[Embedding Matrix] EMB --> HIDDEN HIDDEN --> LMHEAD LMHEAD --> VOCAB EMB -.共享权重.- LMHEAD
十一、CausalLM中的forward
训练时:
logits = self.lm_head(hidden_states)
得到:
[B,L,V]
例如:
[2,128,50000]
然后:
x = logits[:,:-1,:]
y = labels[:,1:]
变成:
输入:
我 爱 中国
预测:
爱 中国 人民
标签:
爱 中国 人民
即:
flowchart LR Token1 --> Predict2 Token2 --> Predict3 Token3 --> Predict4
这就是:
Next Token Prediction
十二、Generate到底怎么工作
这是很多人第一次读源码最容易迷糊的地方。
实际上:
for _ in range(max_new_tokens):
就是不断循环。
流程图:
flowchart TD
A[输入Token]
A --> B[Forward]
B --> C[Logits]
C --> D[TopK TopP]
D --> E[采样Next Token]
E --> F[拼接到输入]
F --> G{EOS?}
G --否--> B
G --是--> H[结束]举例:
用户:
今天
Tokenizer:
[100]
第一次:
[100]
预测:
天气
得到:
[100,200]
第二次:
[100,200]
预测:
很好
得到:
[100,200,300]
第三次:
[100,200,300]
预测:
。
得到:
[100,200,300,400]
最终:
今天天气很好。
十三、KV Cache为什么快
关键代码:
past_key_values
第一次:
Token1
Token2
Token3
计算:
K1 K2 K3
V1 V2 V3
缓存:
KV Cache
下一轮只算:
Token4
生成:
K4
V4
拼接:
[K1 K2 K3 K4]
[V1 V2 V3 V4]
不用重新算前面三个。
图:
flowchart LR T1 --> K1 T2 --> K2 T3 --> K3 K1 --> CACHE K2 --> CACHE K3 --> CACHE NEW[New Token] NEW --> K4 K4 --> CACHE
复杂度从:
下降到接近:
这就是推理加速核心。
一张图总结整个 CausalLM
flowchart TD
A[Input IDs]
A --> B[Embedding]
B --> C[Transformer Block 1]
C --> D[Transformer Block 2]
D --> E[...]
E --> F[Transformer Block N]
F --> G[Final RMSNorm]
G --> H[LM Head]
H --> I[Vocabulary Logits]
I --> J[TopK TopP Sampling]
J --> K[Next Token]
K --> L[Append To Input]
L --> M{EOS?}
M -- No --> B
M -- Yes --> N[Finish]这张图基本对应了完整执行链路,也是 Llama、Qwen、DeepSeek、GPT 等现代自回归大模型的标准架构。